GRU
门控循环单元(Gated Recurrent Unit)。对输入序列按时间步做 GRU 递推,支持单向与双向。
对每个时间步 \(t\),由输入 \(x_t\) 与上一隐藏状态 \(h_{t-1}\) 计算:
其中 \(\sigma\) 为 Sigmoid,\(\odot\) 为逐元素乘。候选门 \(n_t\) 采用 PyTorch 写法(先对 \(h_{(t-1)}\) 做线性再与 \(r_t\) 相乘),与原始论文中先做 \(r_t \odot h_{(t-1)}\) 再乘权重的形式不同。
权重与偏置按门顺序打包为 reset / update / candidate 三段;单向时 num_directions = 1,双向时为 2。
典型形状(单向):
input:
[seq_len, batch, input_size]weight_g:
[num_directions, hidden_size × 3, input_size]weight_r:
[num_directions, hidden_size × 3, hidden_size]input_bias / state_bias:
[num_directions, hidden_size × 3]hidden_state:
[num_directions, batch, hidden_size](调用前写入初值 h0,调用后为最终隐藏状态)output:
[seq_len, batch, num_directions × hidden_size]
- 输入:
input - 输入序列地址
hidden_state - 隐藏状态地址(入出参:初值进、终态出)
buffer - 长度为 4 的二级指针数组,临时工作区(见 note)
gru_param - GruParameter 指针,权重指针与尺寸等,见下
core_mask - 核掩码(仅共享存储版本)
GruParameter 定义:
1typedef struct GruParameter {
2 void *weight_g_; /* 输入-隐藏权重 */
3 void *weight_r_; /* 隐藏-隐藏权重 */
4 void *input_bias_; /* 输入-隐藏偏置 */
5 void *state_bias_; /* 隐藏-隐藏偏置 */
6 int input_size_; /* 输入特征维 */
7 int hidden_size_; /* 隐藏特征维 */
8 int seq_len_; /* 序列长度(可含 padding) */
9 int batch_; /* batch */
10 int output_step_; /* 时间步在 output 上的步长,常为 batch × hidden × num_directions */
11 int bidirectional_; /* 非 0 为双向 */
12 int input_row_align_; /* 保留字段 */
13 int input_col_align_; /* 输入侧列对齐,测试常用 hidden_size */
14 int state_row_align_; /* 保留字段 */
15 int state_col_align_; /* 状态侧列对齐,测试常用 hidden_size */
16 int check_seq_len_; /* 有效序列长度;大于此长度的输出会被清零 */
17} GruParameter;
- 输出:
output - 各时间步隐藏状态序列
hidden_state - 最终隐藏状态(原地更新)
- 支持平台:
FT78NEMT7004
备注
FT78NE 支持 int8、fp32
MT7004 支持 fp16、fp32
权重、偏置通过
gru_param中的指针传入,不作为独立函数形参buffer[1]需能容纳seq_len × batch × hidden_size × 3个元素(门预计算);buffer[3]需能容纳batch × hidden_size × 3个元素(单步状态门);buffer[0]、buffer[2]亦须预先分配
共享存储版本:
-
void i8_gru_s(int8_t *output, int8_t *input, int8_t *hidden_state, int8_t **buffer, GruParameter *gru_param, int core_mask)
-
void hp_gru_s(float16 *output, float16 *input, float16 *hidden_state, float16 **buffer, GruParameter *gru_param, int core_mask)
-
void fp_gru_s(float *output, float *input, float *hidden_state, float **buffer, GruParameter *gru_param, int core_mask)
C调用示例:
1// MT7004 示例(共享存储多核,DDR 地址)
2void TestGruSMCFp32(int core_mask) {
3 int core_id = get_core_id();
4 int logic_core_id = GetLogicCoreId(core_mask, core_id);
5 int core_num = GetCoreNum(core_mask);
6 int seq_len = 4, batch = 2, input_size = 2, hidden_size = 2;
7 int bidirectional = 0, num_directions = 1;
8 float *output = (float *)0x81000000;
9 float *input = (float *)0x82000000;
10 float *weight_g = (float *)0x83000000;
11 float *weight_r = (float *)0x84000000;
12 float *input_bias = (float *)0x84100000;
13 float *state_bias = (float *)0x84200000;
14 float *hidden_state = (float *)0x84300000;
15 float **buffer = (float **)0x84400000;
16 float *output_hidden_state = (float *)0x84500000;
17 GruParameter *param = (GruParameter *)0x84600000;
18 if (logic_core_id == 0) {
19 memcpy(output_hidden_state, hidden_state,
20 (size_t)(num_directions * batch * hidden_size) * sizeof(float));
21 buffer[0] = (float *)0x85000000;
22 buffer[1] = (float *)0x86000000;
23 buffer[2] = (float *)0x86100000;
24 buffer[3] = (float *)0x86200000;
25 param->weight_g_ = weight_g;
26 param->weight_r_ = weight_r;
27 param->input_bias_ = input_bias;
28 param->state_bias_ = state_bias;
29 param->input_size_ = input_size;
30 param->hidden_size_ = hidden_size;
31 param->seq_len_ = seq_len;
32 param->batch_ = batch;
33 param->output_step_ = batch * hidden_size * num_directions;
34 param->bidirectional_ = bidirectional;
35 param->input_col_align_ = hidden_size;
36 param->state_col_align_ = hidden_size;
37 param->check_seq_len_ = seq_len;
38 }
39 sys_bar(0, core_num);
40 fp_gru_s(output, input, output_hidden_state, buffer, param, core_mask);
41}
42
43void main() {
44 int core_mask = 0b1111;
45 TestGruSMCFp32(core_mask);
46}
私有存储版本:
-
void i8_gru_p(int8_t *output, int8_t *input, int8_t *hidden_state, int8_t **buffer, GruParameter *gru_param)
-
void hp_gru_p(float16 *output, float16 *input, float16 *hidden_state, float16 **buffer, GruParameter *gru_param)
-
void fp_gru_p(float *output, float *input, float *hidden_state, float **buffer, GruParameter *gru_param)
C调用示例:
1// MT7004 示例(私有存储单核,AM 地址)
2void TestGruAMFp32(void) {
3 int seq_len = 4, batch = 2, input_size = 2, hidden_size = 2;
4 int bidirectional = 0, num_directions = 1;
5 float *output = (float *)0x10000000;
6 float *input = (float *)0x10010000;
7 float *weight_g = (float *)0x10020000;
8 float *weight_r = (float *)0x10024000;
9 float *input_bias = (float *)0x10028000;
10 float *state_bias = (float *)0x1002C000;
11 float *hidden_state = (float *)0x10030000;
12 float **buffer = (float **)0x10034000;
13 float *output_hidden_state = (float *)0x10038000;
14 GruParameter *param = (GruParameter *)0x1003C000;
15 memcpy(output_hidden_state, hidden_state,
16 (size_t)(num_directions * batch * hidden_size) * sizeof(float));
17 buffer[0] = (float *)0x10048000;
18 buffer[1] = (float *)0x1004C000;
19 buffer[2] = (float *)0x10050000;
20 buffer[3] = (float *)0x10054000;
21 param->weight_g_ = weight_g;
22 param->weight_r_ = weight_r;
23 param->input_bias_ = input_bias;
24 param->state_bias_ = state_bias;
25 param->input_size_ = input_size;
26 param->hidden_size_ = hidden_size;
27 param->seq_len_ = seq_len;
28 param->batch_ = batch;
29 param->output_step_ = batch * hidden_size * num_directions;
30 param->bidirectional_ = bidirectional;
31 param->input_col_align_ = hidden_size;
32 param->state_col_align_ = hidden_size;
33 param->check_seq_len_ = seq_len;
34 fp_gru_p(output, input, output_hidden_state, buffer, param);
35}
36
37void main() {
38 TestGruAMFp32();
39}